Create a System class in Python - #325
EricBoittier wants to merge 5 commits into
Conversation
53e7927 to
bca7450
Compare
Luthaf
left a comment
There was a problem hiding this comment.
Did you check for compatibility with torch.compile/jax.jit?
I generated a notebook to answer these questions (attached as PDF, sorry for LLM-braining It shows:
|
67d8cd2 to
71002d0
Compare
|
I'd say it's fine for now if we can not construct a system inside a |
Luthaf
left a comment
There was a problem hiding this comment.
I do not understand the new requirement of releasing all arrays views before adding pairs/custom data. Can you explain?
The two APIs do not share buffers. We can keep this, or allow adds with live views (interior mutability on pairs / custom_data only). Which do you want? On the topic of mutability, jax arrays strictly immutable. In torch, arrays may still look writable depending on version; the C comment on get_data is still “do not modify” (c_api/system.rs 168–170). |
|
Ah, right there is the
This should already be handled by DLPack, which has a flag for immutable data. So unless we gain something by specifying it ourself, I'd leave read-only considerations to DLPack and not mention it ourself. |
| if module.startswith("torch") and name == "Tensor": | ||
| return "torch" | ||
| if module.startswith("jax") or module.startswith("jaxlib"): | ||
| return "jax" |
There was a problem hiding this comment.
should maybe use isinstance instead since we could have subclasses of these. See https://github.com/metatensor/metatensor/blob/9dfa936675f8ee92087ca63d122ac04f76333447/python/metatensor_operations/metatensor_operations/_dispatch.py#L44
| raise ValueError(f"Unknown arrays backend: {backend}") | ||
|
|
||
|
|
||
| def _pair_options_json(options: PairListOptions) -> bytes: |
| The returned array uses the configured :py:attr:`arrays_backend` and is | ||
| a read-only view. |
There was a problem hiding this comment.
| The returned array uses the configured :py:attr:`arrays_backend` and is | |
| a read-only view. | |
| The returned array uses the configured :py:attr:`arrays_backend`. |
and same below
9b44483 to
15fd62d
Compare
4773782 to
6c05848
Compare
5ec31c2 to
09b9cf2
Compare
68614de to
33fd484
Compare
|
I rebased and squashed all commits into a single one, but there are still some unresolved review comments here. to pull the rebased changes: git checkout python-wrap-system
git fetch
git reset --hard origin/python-wrap-system # (assuming the remote it called origin) |
Co-authored-by: Cursor <cursoragent@cursor.com>
Added comments to clarify the use of `ManuallyDrop` in `mta_system_add_pairs` and `mta_system_add_custom_data` functions, ensuring that the caller's system remains valid even when these functions fail. Introduced a new test case to verify that the system is not dropped on failure during pair and custom data additions.
Introduced functions to manage JAX tracers and static vectors, allowing the System class to handle JAX arrays as part of its internal structure. Updated the class to maintain arrays in Python when JAX tracers are detected, ensuring compatibility with JAX transformations. Added tests to validate JAX integration, including behavior with JAX's jit and value_and_grad functions, and enforced restrictions on pair lists and custom data in JAX contexts.
2247bd6 to
b02b9ff
Compare
|
Rebased onto current The review notes from the last pass are in this branch: backend detection uses After that I registered You said this did not have to land here, as long as Written with an LLM assistant. |
Contributor (creator of pull-request) checklist
Reviewer checklist